#include <stdint.h>

#include <cuda_bf16.h>

#include <cuda_runtime.h>



// xcore1000's CUDA-compatible compiler does not expose NVIDIA's __dp4a.

// This is a correctness-first replacement: each int32 stores four signed

// int8 values in little-endian byte order.

__device__ inline int32_t signed_byte(uint32_t x) {

    x &= 0xffu;

    return (int32_t)(x ^ 0x80u) - 128;

}



__device__ inline int32_t dp4a_compat(int32_t a, int32_t b, int32_t acc) {

    uint32_t ua = (uint32_t)a;

    uint32_t ub = (uint32_t)b;

    acc += signed_byte(ua) * signed_byte(ub);

    acc += signed_byte(ua >> 8) * signed_byte(ub >> 8);

    acc += signed_byte(ua >> 16) * signed_byte(ub >> 16);

    acc += signed_byte(ua >> 24) * signed_byte(ub >> 24);

    return acc;

}



__global__ void w8a8_moe_gemm_kernel(

    const int8_t* __restrict__ a,

    const int8_t* __restrict__ b_col_major,

    const float* __restrict__ scale_a,

    const float* __restrict__ scale_b,

    const float* __restrict__ moe_weights,

    const int32_t* __restrict__ token_ids,

    const int32_t* __restrict__ expert_ids,

    int K, int N, int topk,

    __nv_bfloat16* __restrict__ out)

{

    int n_base = blockIdx.x * 128;

    int m_base = blockIdx.y * 128;

    int expert = expert_ids[blockIdx.y];



    int tid = threadIdx.x;

    int warp_id = tid / 32;

    int lane_id = tid & 31;



    int warp_y = warp_id / 2;

    int warp_x = warp_id & 1;

    int my = lane_id / 8;

    int mx = lane_id & 7;



    int m_idx[8];

    int n_idx[8];

#pragma unroll

    for (int i = 0; i < 8; ++i) {

        m_idx[i] = warp_y * 32 + my + i * 4;

    }

#pragma unroll

    for (int j = 0; j < 8; ++j) {

        n_idx[j] = warp_x * 64 + mx + j * 8;

    }



    __shared__ int32_t smem_A[2][128 * 17];

    __shared__ int32_t smem_B[2][128 * 17];



    int32_t accum[8][8] = {0};



#pragma unroll

    for (int step = 0; step < 2; ++step) {

        int load_idx = step * 256 + tid;

        int row = load_idx / 4;

        int col_int4 = load_idx & 3;



        int r = m_base + row;

        int token = token_ids[r] / topk;

        int64_t a_idx = (int64_t)token * K;

        int4 va = ((const int4*)(a + a_idx))[col_int4];



        int sa = row * 17 + col_int4 * 4;

        smem_A[0][sa + 0] = va.x;

        smem_A[0][sa + 1] = va.y;

        smem_A[0][sa + 2] = va.z;

        smem_A[0][sa + 3] = va.w;



        int64_t b_idx = (int64_t)expert * N * K + (int64_t)(n_base + row) * K;

        int4 vb = ((const int4*)(b_col_major + b_idx))[col_int4];



        int sb = row * 17 + col_int4 * 4;

        smem_B[0][sb + 0] = vb.x;

        smem_B[0][sb + 1] = vb.y;

        smem_B[0][sb + 2] = vb.z;

        smem_B[0][sb + 3] = vb.w;

    }

    __syncthreads();



    for (int k_outer = 0; k_outer < K; k_outer += 64) {

        int comp_buf = (k_outer / 64) & 1;

        int load_buf = 1 - comp_buf;

        int next_k = k_outer + 64;



        if (next_k < K) {

#pragma unroll

            for (int step = 0; step < 2; ++step) {

                int load_idx = step * 256 + tid;

                int row = load_idx / 4;

                int col_int4 = load_idx & 3;



                int r = m_base + row;

                int token = token_ids[r] / topk;

                int64_t a_idx = (int64_t)token * K + next_k;

                int4 va = ((const int4*)(a + a_idx))[col_int4];



                int sa = row * 17 + col_int4 * 4;

                smem_A[load_buf][sa + 0] = va.x;

                smem_A[load_buf][sa + 1] = va.y;

                smem_A[load_buf][sa + 2] = va.z;

                smem_A[load_buf][sa + 3] = va.w;



                int64_t b_idx = (int64_t)expert * N * K +

                                (int64_t)(n_base + row) * K + next_k;

                int4 vb = ((const int4*)(b_col_major + b_idx))[col_int4];



                int sb = row * 17 + col_int4 * 4;

                smem_B[load_buf][sb + 0] = vb.x;

                smem_B[load_buf][sb + 1] = vb.y;

                smem_B[load_buf][sb + 2] = vb.z;

                smem_B[load_buf][sb + 3] = vb.w;

            }

        }



#pragma unroll

        for (int k_step = 0; k_step < 16; ++k_step) {

            int32_t reg_A[8];

            int32_t reg_B[8];



#pragma unroll

            for (int i = 0; i < 8; ++i) {

                reg_A[i] = smem_A[comp_buf][m_idx[i] * 17 + k_step];

            }

#pragma unroll

            for (int j = 0; j < 8; ++j) {

                reg_B[j] = smem_B[comp_buf][n_idx[j] * 17 + k_step];

            }



#pragma unroll

            for (int i = 0; i < 8; ++i) {

#pragma unroll

                for (int j = 0; j < 8; ++j) {

                    accum[i][j] = dp4a_compat(reg_A[i], reg_B[j], accum[i][j]);

                }

            }

        }

        __syncthreads();

    }



    float scale_row[8];

#pragma unroll

    for (int i = 0; i < 8; ++i) {

        int r = m_base + m_idx[i];

        int token = token_ids[r] / topk;

        scale_row[i] = scale_a[token] * moe_weights[r];

    }



    float scale_col[8];

#pragma unroll

    for (int j = 0; j < 8; ++j) {

        int n = n_base + n_idx[j];

        scale_col[j] = scale_b[(int64_t)expert * N + n];

    }



#pragma unroll

    for (int i = 0; i < 8; ++i) {

        int r = m_base + m_idx[i];

#pragma unroll

        for (int j = 0; j < 8; ++j) {

            int n = n_base + n_idx[j];

            float v = (float)accum[i][j] * scale_row[i] * scale_col[j];

            out[(int64_t)r * N + n] = __float2bfloat16(v);

        }

    }

}



static size_t device_allocation_size(const void* p) {

    mcDrvDeviceptr_t base = 0;

    size_t size = 0;

    (void)wcuMemGetAddressRange(&base, &size, (mcDrvDeviceptr_t)(uintptr_t)p);

    return size;

}



extern "C" void run_kernel(

    const int8_t* a,

    const int8_t* b_col_major,

    const float* scale_a,

    const float* scale_b,

    const float* moe_weights,

    const int32_t* token_ids,

    const int32_t* expert_ids,

    int64_t topk,

    __nv_bfloat16* out)

{

    size_t b_size = device_allocation_size(b_col_major);

    size_t out_size = device_allocation_size(out);



    int N = 7168;

    int K = 2048;

    if (b_size > 5000000000ULL) {

        N = 4096;

        K = 7168;

    }



    int EM = 4096;

    if (out_size > 128ULL * 1024ULL * 1024ULL) {

        EM = 32768;

    } else if (out_size == 0) {

        // Last-resort fallback if allocation-size probing is unavailable.

        int32_t host_tokens[4096];

        cudaMemcpy(host_tokens, token_ids, sizeof(host_tokens), cudaMemcpyDeviceToHost);

        int max_token_id = 0;

        for (int i = 0; i < 4096; ++i) {

            if (host_tokens[i] > max_token_id) {

                max_token_id = host_tokens[i];

            }

        }

        if (max_token_id >= 4096) {

            EM = 32768;

        }

    }



    dim3 block(256);

    dim3 grid(N / 128, EM / 128);

    w8a8_moe_gemm_kernel<<<grid, block>>>(

        a, b_col_major, scale_a, scale_b, moe_weights,

        token_ids, expert_ids, K, N, (int)topk, out);

}